#!/usr/bin/env python3
"""청구항변반박전략서.md와 evidence_indexed.json에서 complaint_information.md를 생성한다."""

from __future__ import annotations

import itertools
import json
import re
import sys
import unicodedata
from typing import Any

import httpx


# =============================================================================
# MCP / constants
# =============================================================================
LOCALDOCS_URL = "http://mcp-localdocs:8012/mcp"
MCP_HEADERS = {
    "Content-Type": "application/json",
    "Accept": "application/json, text/event-stream",
}
CLIENT = httpx.Client(timeout=60)
MSG_ID_COUNTER = itertools.count(10)

SOURCE_DOC_NAME = unicodedata.normalize("NFC", "청구항변반박전략서.md")
EVIDENCE_DOC_NAME = unicodedata.normalize("NFC", "evidence_indexed.json")
OUTPUT_DOC_NAME = "complaint_information.md"


def next_msg_id() -> int:
    return next(MSG_ID_COUNTER)


def _init_mcp_session() -> None:
    response = CLIENT.post(
        LOCALDOCS_URL,
        json={
            "jsonrpc": "2.0",
            "id": 1,
            "method": "initialize",
            "params": {
                "protocolVersion": "2025-03-26",
                "capabilities": {},
                "clientInfo": {"name": "stage5-complaint-info", "version": "1.1"},
            },
        },
        headers=MCP_HEADERS,
    )
    response.raise_for_status()

    session_id = response.headers.get("mcp-session-id")
    if session_id:
        MCP_HEADERS["mcp-session-id"] = session_id

    CLIENT.post(
        LOCALDOCS_URL,
        json={"jsonrpc": "2.0", "method": "notifications/initialized"},
        headers=MCP_HEADERS,
    ).raise_for_status()


def _parse_mcp(text: str) -> dict[str, Any] | None:
    for line in text.strip().split("\n"):
        if line.startswith("data: "):
            try:
                return json.loads(line[6:])
            except json.JSONDecodeError:
                return None

    try:
        return json.loads(text)
    except json.JSONDecodeError:
        return None


def _call_tool(name: str, arguments: dict[str, Any], msg_id: int) -> dict[str, Any]:
    response = CLIENT.post(
        LOCALDOCS_URL,
        json={
            "jsonrpc": "2.0",
            "id": msg_id,
            "method": "tools/call",
            "params": {"name": name, "arguments": arguments},
        },
        headers=MCP_HEADERS,
    )
    response.raise_for_status()

    parsed = _parse_mcp(response.text)
    if not parsed or "result" not in parsed:
        raise RuntimeError(f"MCP {name} failed: {response.text[:300]}")
    return parsed


def _unwrap_text(parsed: dict[str, Any], doc_name: str) -> str:
    text = (parsed["result"].get("content") or [{}])[0].get("text", "")
    if not text:
        raise RuntimeError(f"Empty response for {doc_name}")

    try:
        outer = json.loads(text)
    except json.JSONDecodeError:
        return text

    if isinstance(outer, dict) and "results" in outer:
        results = outer.get("results") or []
        if not results:
            raise RuntimeError(f"No results for {doc_name}")

        inner = results[0].get("content") or results[0].get("text") or ""
        if not inner:
            raise RuntimeError(f"Empty content for {doc_name}")
        return inner if isinstance(inner, str) else json.dumps(inner, ensure_ascii=False)

    return json.dumps(outer, ensure_ascii=False) if not isinstance(outer, str) else outer


def read_text(doc_name: str) -> str:
    parsed = _call_tool("read_docs", {"doc_names": [doc_name]}, next_msg_id())
    return _unwrap_text(parsed, doc_name)


def write_doc(path: str, content: str) -> None:
    _call_tool("write_file", {"path": path, "content": content}, next_msg_id())


# =============================================================================
# Document extraction logic
# =============================================================================
MASTER_SECTION_HEADER = "## 3. 청구 마스터 표"
OVERVIEW_HEADER = "### 3.1 청구권 개요"
DETAIL_HEADER = "### 3.2 개별 청구권 상세 정보"
NEXT_TOP_LEVEL_HEADER_PATTERN = re.compile(r"^##\s+\d+\.", re.MULTILINE)
DETAIL_BLOCK_PATTERN = re.compile(
    r"(?ms)^(####\s+C-\d{3}:.*?\n\n<div class=\"full-width-table-wrap\".*?</div>)\s*"
)
EVIDENCE_INDEX_PATTERN = re.compile(r"^E-(\d+)$")


def find_section_bounds(text: str, start_header: str, end_header: str | None = None) -> tuple[int, int]:
    start = text.find(start_header)
    if start == -1:
        raise RuntimeError(f"헤더를 찾을 수 없습니다: {start_header}")

    section_start = start + len(start_header)
    if end_header:
        end = text.find(end_header, section_start)
        if end == -1:
            raise RuntimeError(f"종료 헤더를 찾을 수 없습니다: {end_header}")
        return section_start, end

    match = NEXT_TOP_LEVEL_HEADER_PATTERN.search(text, section_start)
    if not match:
        return section_start, len(text)
    return section_start, match.start()


def extract_first_table_block(section_text: str, context_name: str) -> str:
    div_start = section_text.find('<div class="full-width-table-wrap"')
    if div_start == -1:
        raise RuntimeError(f"{context_name}에서 표 블록 시작을 찾을 수 없습니다.")

    div_end = section_text.find("\n</div>", div_start)
    if div_end == -1:
        raise RuntimeError(f"{context_name}에서 표 블록 종료를 찾을 수 없습니다.")

    return section_text[div_start : div_end + len("\n</div>")].strip()


def extract_overview_table(master_section_text: str) -> str:
    start, end = find_section_bounds(master_section_text, OVERVIEW_HEADER, DETAIL_HEADER)
    overview_section = master_section_text[start:end]
    return extract_first_table_block(overview_section, "청구권 개요")


def extract_detail_blocks(master_section_text: str) -> list[str]:
    start, end = find_section_bounds(master_section_text, DETAIL_HEADER)
    detail_section = master_section_text[start:end]
    blocks = [match.group(1).strip() for match in DETAIL_BLOCK_PATTERN.finditer(detail_section)]
    if not blocks:
        raise RuntimeError("개별 청구권 상세 정보 블록을 찾을 수 없습니다.")
    return blocks


def parse_evidence_json(raw_text: str) -> list[dict[str, Any]]:
    try:
        parsed = json.loads(raw_text.strip())
    except json.JSONDecodeError as exc:
        raise RuntimeError("evidence_indexed.json을 JSON으로 파싱할 수 없습니다.") from exc

    if not isinstance(parsed, list):
        raise RuntimeError("evidence_indexed.json은 객체 배열이어야 합니다.")

    return [item for item in parsed if isinstance(item, dict)]


def _clean_string(value: Any) -> str:
    return value.strip() if isinstance(value, str) else ""


def build_evidence_lines(evidence_items: list[dict[str, Any]]) -> list[str]:
    deduped_entries: dict[str, tuple[int, str]] = {}

    for item in evidence_items:
        evidence_index = _clean_string(item.get("evidence_index"))
        title_normalized = _clean_string(item.get("title_normalized"))

        if not title_normalized:
            continue

        match = EVIDENCE_INDEX_PATTERN.fullmatch(evidence_index)
        if not match:
            continue

        if evidence_index in deduped_entries:
            continue

        evidence_number = int(match.group(1))
        deduped_entries[evidence_index] = (evidence_number, title_normalized)

    sorted_entries = sorted(deduped_entries.values(), key=lambda entry: entry[0])
    return [
        f"갑 제{evidence_number}호증 - {title_normalized}"
        for evidence_number, title_normalized in sorted_entries
    ]


def assemble_document(overview_table: str, detail_blocks: list[str], evidence_lines: list[str]) -> str:
    parts = [
        "# 청구상세정보",
        "",
        "## 1. 청구권 개요",
        "",
        overview_table,
        "",
        "## 2. 개별 청구권 상세 정보",
        "",
        "\n\n".join(detail_blocks),
        "",
        "## 3. 서증 목록",
        "",
    ]

    parts.extend(evidence_lines)
    parts.append("")
    return "\n".join(parts)


# =============================================================================
# Main
# =============================================================================
def main() -> None:
    _init_mcp_session()

    source_text = read_text(SOURCE_DOC_NAME)
    evidence_text = read_text(EVIDENCE_DOC_NAME)
    print(f"[stage5] read {SOURCE_DOC_NAME} ({len(source_text)} chars)", file=sys.stderr)
    print(f"[stage5] read {EVIDENCE_DOC_NAME} ({len(evidence_text)} chars)", file=sys.stderr)

    master_start, master_end = find_section_bounds(source_text, MASTER_SECTION_HEADER)
    master_section_text = source_text[master_start:master_end]

    overview_table = extract_overview_table(master_section_text)
    detail_blocks = extract_detail_blocks(master_section_text)
    evidence_items = parse_evidence_json(evidence_text)
    evidence_lines = build_evidence_lines(evidence_items)
    output_text = assemble_document(overview_table, detail_blocks, evidence_lines)

    write_doc(OUTPUT_DOC_NAME, output_text)
    print(
        f"[stage5] saved: {OUTPUT_DOC_NAME} "
        f"(overview=1, detail_blocks={len(detail_blocks)}, evidence_lines={len(evidence_lines)})",
        file=sys.stderr,
    )
    print(
        json.dumps(
            {
                "status": "ok",
                "detail_blocks": len(detail_blocks),
                "evidence_lines": len(evidence_lines),
            }
        )
    )


if __name__ == "__main__":
    main()
